GRU

门控循环单元(Gated Recurrent Unit)。对输入序列按时间步做 GRU 递推,支持单向与双向。

对每个时间步 \(t\),由输入 \(x_t\) 与上一隐藏状态 \(h_{t-1}\) 计算:

\begin{align*} r_t &= \sigma(W_{ir}x_t + b_{ir} + W_{hr}h_{(t-1)} + b_{hr}) \\ z_t &= \sigma(W_{iz}x_t + b_{iz} + W_{hz}h_{(t-1)} + b_{hz}) \\ n_t &= \tanh(W_{in}x_t + b_{in} + r_t \odot (W_{hn}h_{(t-1)} + b_{hn})) \\ h_t &= (1 - z_t) \odot n_t + z_t \odot h_{(t-1)} \end{align*}

其中 \(\sigma\) 为 Sigmoid,\(\odot\) 为逐元素乘。候选门 \(n_t\) 采用 PyTorch 写法(先对 \(h_{(t-1)}\) 做线性再与 \(r_t\) 相乘),与原始论文中先做 \(r_t \odot h_{(t-1)}\) 再乘权重的形式不同。

权重与偏置按门顺序打包为 reset / update / candidate 三段;单向时 num_directions = 1,双向时为 2。

典型形状(单向):

  • input: [seq_len, batch, input_size]

  • weight_g: [num_directions, hidden_size × 3, input_size]

  • weight_r: [num_directions, hidden_size × 3, hidden_size]

  • input_bias / state_bias: [num_directions, hidden_size × 3]

  • hidden_state: [num_directions, batch, hidden_size] (调用前写入初值 h0,调用后为最终隐藏状态)

  • output: [seq_len, batch, num_directions × hidden_size]

输入:
  • input - 输入序列地址

  • hidden_state - 隐藏状态地址(入出参:初值进、终态出)

  • buffer - 长度为 4 的二级指针数组,临时工作区(见 note)

  • gru_param - GruParameter 指针,权重指针与尺寸等,见下

  • core_mask - 核掩码(仅共享存储版本)

GruParameter 定义:

 1typedef struct GruParameter {
 2    void *weight_g_;       /* 输入-隐藏权重 */
 3    void *weight_r_;       /* 隐藏-隐藏权重 */
 4    void *input_bias_;     /* 输入-隐藏偏置 */
 5    void *state_bias_;     /* 隐藏-隐藏偏置 */
 6    int input_size_;       /* 输入特征维 */
 7    int hidden_size_;      /* 隐藏特征维 */
 8    int seq_len_;          /* 序列长度(可含 padding) */
 9    int batch_;            /* batch */
10    int output_step_;      /* 时间步在 output 上的步长,常为 batch × hidden × num_directions */
11    int bidirectional_;    /* 非 0 为双向 */
12    int input_row_align_;  /* 保留字段 */
13    int input_col_align_;  /* 输入侧列对齐,测试常用 hidden_size */
14    int state_row_align_;  /* 保留字段 */
15    int state_col_align_;  /* 状态侧列对齐,测试常用 hidden_size */
16    int check_seq_len_;    /* 有效序列长度;大于此长度的输出会被清零 */
17} GruParameter;
输出:
  • output - 各时间步隐藏状态序列

  • hidden_state - 最终隐藏状态(原地更新)

支持平台:

FT78NE MT7004

备注

  • FT78NE 支持 int8、fp32

  • MT7004 支持 fp16、fp32

  • 权重、偏置通过 gru_param 中的指针传入,不作为独立函数形参

  • buffer[1] 需能容纳 seq_len × batch × hidden_size × 3 个元素(门预计算);buffer[3] 需能容纳 batch × hidden_size × 3 个元素(单步状态门);buffer[0]buffer[2] 亦须预先分配

共享存储版本:

void i8_gru_s(int8_t *output, int8_t *input, int8_t *hidden_state, int8_t **buffer, GruParameter *gru_param, int core_mask)
void hp_gru_s(float16 *output, float16 *input, float16 *hidden_state, float16 **buffer, GruParameter *gru_param, int core_mask)
void fp_gru_s(float *output, float *input, float *hidden_state, float **buffer, GruParameter *gru_param, int core_mask)

C调用示例:

 1// MT7004 示例(共享存储多核,DDR 地址)
 2void TestGruSMCFp32(int core_mask) {
 3    int core_id = get_core_id();
 4    int logic_core_id = GetLogicCoreId(core_mask, core_id);
 5    int core_num = GetCoreNum(core_mask);
 6    int seq_len = 4, batch = 2, input_size = 2, hidden_size = 2;
 7    int bidirectional = 0, num_directions = 1;
 8    float *output = (float *)0x81000000;
 9    float *input = (float *)0x82000000;
10    float *weight_g = (float *)0x83000000;
11    float *weight_r = (float *)0x84000000;
12    float *input_bias = (float *)0x84100000;
13    float *state_bias = (float *)0x84200000;
14    float *hidden_state = (float *)0x84300000;
15    float **buffer = (float **)0x84400000;
16    float *output_hidden_state = (float *)0x84500000;
17    GruParameter *param = (GruParameter *)0x84600000;
18    if (logic_core_id == 0) {
19        memcpy(output_hidden_state, hidden_state,
20               (size_t)(num_directions * batch * hidden_size) * sizeof(float));
21        buffer[0] = (float *)0x85000000;
22        buffer[1] = (float *)0x86000000;
23        buffer[2] = (float *)0x86100000;
24        buffer[3] = (float *)0x86200000;
25        param->weight_g_ = weight_g;
26        param->weight_r_ = weight_r;
27        param->input_bias_ = input_bias;
28        param->state_bias_ = state_bias;
29        param->input_size_ = input_size;
30        param->hidden_size_ = hidden_size;
31        param->seq_len_ = seq_len;
32        param->batch_ = batch;
33        param->output_step_ = batch * hidden_size * num_directions;
34        param->bidirectional_ = bidirectional;
35        param->input_col_align_ = hidden_size;
36        param->state_col_align_ = hidden_size;
37        param->check_seq_len_ = seq_len;
38    }
39    sys_bar(0, core_num);
40    fp_gru_s(output, input, output_hidden_state, buffer, param, core_mask);
41}
42
43void main() {
44    int core_mask = 0b1111;
45    TestGruSMCFp32(core_mask);
46}

私有存储版本:

void i8_gru_p(int8_t *output, int8_t *input, int8_t *hidden_state, int8_t **buffer, GruParameter *gru_param)
void hp_gru_p(float16 *output, float16 *input, float16 *hidden_state, float16 **buffer, GruParameter *gru_param)
void fp_gru_p(float *output, float *input, float *hidden_state, float **buffer, GruParameter *gru_param)

C调用示例:

 1// MT7004 示例(私有存储单核,AM 地址)
 2void TestGruAMFp32(void) {
 3    int seq_len = 4, batch = 2, input_size = 2, hidden_size = 2;
 4    int bidirectional = 0, num_directions = 1;
 5    float *output = (float *)0x10000000;
 6    float *input = (float *)0x10010000;
 7    float *weight_g = (float *)0x10020000;
 8    float *weight_r = (float *)0x10024000;
 9    float *input_bias = (float *)0x10028000;
10    float *state_bias = (float *)0x1002C000;
11    float *hidden_state = (float *)0x10030000;
12    float **buffer = (float **)0x10034000;
13    float *output_hidden_state = (float *)0x10038000;
14    GruParameter *param = (GruParameter *)0x1003C000;
15    memcpy(output_hidden_state, hidden_state,
16           (size_t)(num_directions * batch * hidden_size) * sizeof(float));
17    buffer[0] = (float *)0x10048000;
18    buffer[1] = (float *)0x1004C000;
19    buffer[2] = (float *)0x10050000;
20    buffer[3] = (float *)0x10054000;
21    param->weight_g_ = weight_g;
22    param->weight_r_ = weight_r;
23    param->input_bias_ = input_bias;
24    param->state_bias_ = state_bias;
25    param->input_size_ = input_size;
26    param->hidden_size_ = hidden_size;
27    param->seq_len_ = seq_len;
28    param->batch_ = batch;
29    param->output_step_ = batch * hidden_size * num_directions;
30    param->bidirectional_ = bidirectional;
31    param->input_col_align_ = hidden_size;
32    param->state_col_align_ = hidden_size;
33    param->check_seq_len_ = seq_len;
34    fp_gru_p(output, input, output_hidden_state, buffer, param);
35}
36
37void main() {
38    TestGruAMFp32();
39}